#!/usr/bin/env python3 """ json_validation.py — JSON 格式校验 + 打分-重试循环 + Format Error Library 为知微股票分析 Pipeline 提供三层校验: 1. JSON 语法解析 (json.loads) 2. JSON Schema 字段匹配(可选) 3. 语义完整性检查(关键字段非空、数值范围合理) 核心功能: - validate_json_syntax(raw) → (score, errors, parsed_or_None) - scored_retry_loop(llm_fn, max_retries=3, threshold=90, schema=None) - FormatErrorLibrary — 持久化错误记录,每日/周汇总 - best_effort_fallback(raw) → dict(尽力恢复) 用法: from json_validation import scored_retry_loop, FormatErrorLibrary def my_llm_call(prompt): # 调用 LLM 返回原始文本 return llm_response result = scored_retry_loop( llm_fn=my_llm_call, prompt=initial_prompt, context={"code": "000001", "name": "平安银行"}, schema=SCHEMA, max_retries=3, threshold=90, ) # result = {"ok": True, "data": {...}, "attempts": 2, "score": 95} 或 # result = {"ok": False, "data": {...}, "attempts": 3, "score": 60, "fallback": True} """ import json import os import re import sys import time from datetime import datetime, timedelta from pathlib import Path from collections import defaultdict # ─── 常量 ─────────────────────────────────────────────────── DEFAULT_THRESHOLD = 90 # 验收阈值(0-100) DEFAULT_MAX_RETRIES = 3 # 最大重试次数 ERROR_LIBRARY_PATH = os.path.join( os.path.dirname(os.path.abspath(__file__)), "data", "format_error_library.json" ) # ─── 标准 JSON Schema 种子字段 ───────────────────────────── # 知微股票分析输出期望包含的字段 STOCK_ANALYSIS_SCHEMA = { "required": [ "code", "name", "action", "stop_loss", "take_profit", "entry_low", "entry_high", "timing_signal", "rr_ratio", ], "optional": [ "price", "cost", "shares", "currency", "tech_snapshot", "signal_factors", "sector_context", "multi_tf_context", "status", "quality_check", ], "types": { "code": str, "name": str, "action": str, "stop_loss": (int, float), "take_profit": (int, float), "entry_low": (int, float), "entry_high": (int, float), "timing_signal": str, "rr_ratio": (int, float), } } # ═══════════════════════════════════════════════════════════ # 第一层:JSON 语法校验 + 打分 # ═══════════════════════════════════════════════════════════ def validate_json_syntax(raw_text): """JSON 语法校验 + 打分(0-100) 返回: (score, errors, parsed_or_None) score: 0-100 整数 errors: 错误描述列表 parsed_or_None: 成功 parse 的 dict 或 None """ errors = [] if not raw_text or not raw_text.strip(): return (0, ["空文本"], None) text = raw_text.strip() # 步骤1: 尝试提取可能的 JSON(从 markdown 代码块中) extracted = _extract_json_block(text) if extracted != text: text = extracted # 步骤2: 初步语法检查 score = 50 # 基础分 # 检查是否有 JSON 的基本特征 if text.startswith("{") or text.startswith("["): score += 10 else: errors.append("不以 { 或 [ 开头,可能不是 JSON") if text.endswith("}") or text.endswith("]"): score += 10 else: errors.append("不以 } 或 ] 结尾,可能被截断") # 步骤3: 尝试解析 try: parsed = json.loads(text) score += 30 # 成功解析加 30 分 # 检查深层嵌套完整性 _check_nested_integrity(parsed, errors) # 最终得分 if not errors: score = min(100, score + 10) # 完美通过再加 10 return (min(100, score), errors, parsed) except json.JSONDecodeError as e: errors.append(f"JSON 语法错误: {e}") # 常见错误类型归类 err_msg = str(e) if "trailing comma" in err_msg.lower() or "extra data" in err_msg.lower(): errors.append("类型: 尾随逗号或多余数据") score -= 5 elif "unterminated" in err_msg.lower(): errors.append("类型: 未闭合的引号/括号") score -= 10 elif "unexpected" in err_msg.lower(): errors.append("类型: 意外的字符/标记") score -= 8 elif "control character" in err_msg.lower(): errors.append("类型: 字符串包含未转义的控制字符") score -= 5 elif "expect value" in err_msg.lower(): errors.append("类型: 缺少值或字段名") score -= 8 # 尝试自动修复 fixed, fixed_desc = _auto_fix_json(text) if fixed: try: parsed = json.loads(fixed) score = max(score, 75) # 修复成功给 75 分 errors.append(f"自动修复成功: {fixed_desc}") return (score, errors, parsed) except json.JSONDecodeError: pass return (max(0, score - 20), errors, None) def _extract_json_block(text): """从文本中提取 JSON 代码块(```json ... ```)""" # 模式1: ```json ... ``` 或 ``` ... ``` m = re.search(r'```(?:json)?\s*(\{.*?\}|\[.*?\])\s*```', text, re.DOTALL) if m: return m.group(1).strip() # 模式2: ... m = re.search(r'\s*(\{.*?\}|\[.*?\])\s*', text, re.DOTALL) if m: return m.group(1).strip() return text def _check_nested_integrity(obj, errors, path=""): """递归检查嵌套结构的完整性""" if isinstance(obj, dict): for k, v in obj.items(): child_path = f"{path}.{k}" if path else k if v is None: errors.append(f"字段 {child_path} 为 null") elif isinstance(v, str) and len(v) > 1000: errors.append(f"字段 {child_path} 字符串过长 ({len(v)} 字符)") else: _check_nested_integrity(v, errors, child_path) elif isinstance(obj, list): if len(obj) > 100: errors.append(f"数组 {path} 过长 ({len(obj)} 项)") for i, item in enumerate(obj): _check_nested_integrity(item, errors, f"{path}[{i}]") def _auto_fix_json(text): """尝试自动修复常见 JSON 格式问题 返回: (fixed_text_or_None, description_or_None) """ original = text fixes = [] # 修复1: 去掉多余尾部括号(可能有多余的 }}}) fixed = text.rstrip() # 计算开括号和闭括号的数量 open_braces = fixed.count('{') close_braces = fixed.count('}') open_brackets = fixed.count('[') close_brackets = fixed.count(']') if close_braces > open_braces: # 去掉多余的 } excess = close_braces - open_braces for _ in range(excess): last_idx = fixed.rfind('}') if last_idx >= 0: fixed = fixed[:last_idx] + fixed[last_idx+1:] fixes.append(f"去掉 {excess} 个多余右括号") if close_brackets > open_brackets: excess = close_brackets - open_brackets for _ in range(excess): last_idx = fixed.rfind(']') if last_idx >= 0: fixed = fixed[:last_idx] + fixed[last_idx+1:] fixes.append(f"去掉 {excess} 个多余右中括号") # 修复2: 字符串内未转义的换行符 result = [] in_str = False for ch in fixed: if ch == '"': in_str = not in_str result.append(ch) elif in_str and ch in '\n\r': result.append('\\n') else: result.append(ch) fixed2 = ''.join(result) if fixed2 != fixed: fixes.append("转义字符串内换行符") fixed = fixed2 # 修复3: 尾随逗号(在 } 或 ] 前的逗号) fixed = re.sub(r',\s*}', '}', fixed) fixed = re.sub(r',\s*\]', ']', fixed) if fixed != original: fixes.append("移除尾随逗号") # 修复4: 如果还不行,尝试补全为有效的对象包装 if not fixed.startswith('{') and not fixed.startswith('['): # 可能是裸值 → 包装为对象 try: json.loads(fixed) except json.JSONDecodeError: pass # 保留原始 if fixed != original: return fixed, '; '.join(fixes) # 无有效修复 return None, None # ═══════════════════════════════════════════════════════════ # 第二层:JSON Schema 校验 # ═══════════════════════════════════════════════════════════ def validate_json_schema(parsed, schema=None): """JSON Schema 字段匹配验证 schema 格式: { "required": ["field1", "field2"], "optional": ["field3"], "types": {"field1": str, "field2": (int, float)}, "min_values": {"rr_ratio": 0}, "non_empty": ["action", "timing_signal"], } 返回: (score, errors) score: 0-100 errors: 字段级别的错误列表 """ if schema is None: schema = STOCK_ANALYSIS_SCHEMA if not isinstance(parsed, dict): return (30, ["顶层不是对象(dict),无法做 Schema 校验"]) errors = [] score = 50 # 基础分 required = schema.get("required", []) optional = schema.get("optional", []) types = schema.get("types", {}) min_values = schema.get("min_values", {}) non_empty = schema.get("non_empty", []) ranges = schema.get("ranges", {}) # {"field": (min, max)} # 检查必要字段 missing_required = [] for field in required: if field not in parsed: missing_required.append(field) if missing_required: errors.append(f"缺少必要字段: {', '.join(missing_required)}") score -= len(missing_required) * 8 else: score += 15 # 检查字段类型 type_errors = [] for field, expected_type in types.items(): if field in parsed and parsed[field] is not None: val = parsed[field] if not isinstance(val, expected_type): type_errors.append(f"{field}: 期望 {expected_type.__name__}, 实际 {type(val).__name__}") if type_errors: errors.extend(type_errors) score -= len(type_errors) * 5 # 检查最小值约束 for field, min_val in min_values.items(): if field in parsed and parsed[field] is not None: try: if float(parsed[field]) < min_val: errors.append(f"{field}={parsed[field]} < 最小值 {min_val}") score -= 5 except (TypeError, ValueError): pass # 检查非空约束 for field in non_empty: if field in parsed: val = parsed[field] if val is None or (isinstance(val, str) and not val.strip()) or val == 0: errors.append(f"{field} 为空/零") score -= 5 # 检查数值范围 for field, (lo, hi) in ranges.items(): if field in parsed and parsed[field] is not None: try: v = float(parsed[field]) if v < lo or v > hi: errors.append(f"{field}={v} 超出合理范围 [{lo}, {hi}]") score -= 3 except (TypeError, ValueError): pass # 未知字段不扣分(可接受额外信息) # 标准化 score score = max(0, min(100, score)) return (score, errors) # ═══════════════════════════════════════════════════════════ # 第三层:打分-重试循环 # ═══════════════════════════════════════════════════════════ def scored_retry_loop(llm_fn, prompt, context=None, schema=None, max_retries=3, threshold=90, error_lib=None, task_id=None): """打分-重试循环 参数: llm_fn: callable(prompt) → raw_text 每次重试时传入附加上一轮反馈的增强 prompt prompt: 初始 prompt 文本 context: dict,额外的上下文信息,用于记录 schema: JSON Schema 定义(可选) max_retries: 最大重试次数(默认 3) threshold: 验收阈值 0-100(默认 90) error_lib: FormatErrorLibrary 实例(可选),用于记录错误 task_id: 任务标识符(可选),用于错误库跟踪 返回: { "ok": True/False, # True=在阈值内通过, False=最终fallback "data": dict or None, # 解析后的 JSON 数据 "raw": str, # 最终原始输出 "attempts": int, # 实际尝试次数 "scores": [int, ...], # 每次尝试的分数 "errors": [str, ...], # 最终错误列表 "fallback": bool, # 是否走了一段 best_effort } """ # 验证最小报价字段 _validate_min_values(schema) best_score = 0 best_data = None best_raw = "" # 构建基础 prompt 系统反馈部分 base_prompt = prompt feedback_history = [] scores = [] for attempt in range(1, max_retries + 1): # 构建本次的 prompt(附加上一轮的格式反馈) current_prompt = base_prompt if feedback_history: feedback_section = "\n\n[Format Feedback from previous attempt]\n" + \ "\n".join(feedback_history) current_prompt = base_prompt + feedback_section # 调用 LLM try: raw = llm_fn(current_prompt) except Exception as e: scores.append(0) feedback_history.append(f"Attempt {attempt}: LLM call failed: {e}") continue if not raw or not raw.strip(): scores.append(0) feedback_history.append(f"Attempt {attempt}: 空响应") continue # 步骤1: 语法校验 syntax_score, syntax_errors, parsed = validate_json_syntax(raw) # 步骤2: Schema 校验(如果语法通过且有 schema) schema_score = 100 schema_errors = [] if parsed is not None and schema: schema_score, schema_errors = validate_json_schema(parsed, schema) # 综合打分: 语法占 60%, Schema 占 40% if parsed is not None: combined_score = int(syntax_score * 0.6 + schema_score * 0.4) else: combined_score = syntax_score # 语法失败时只用语法分 scores.append(combined_score) # 记录错误信息 all_errors = syntax_errors + schema_errors # 记录到 Format Error Library if error_lib and task_id: error_lib.record( task_id=task_id, attempt=attempt, score=combined_score, errors=all_errors, context=context, ) # 追踪最佳结果 if combined_score > best_score: best_score = combined_score best_data = parsed best_raw = raw # 检查是否通过阈值 if combined_score >= threshold: return { "ok": True, "data": parsed, "raw": raw, "attempts": attempt, "scores": scores, "errors": all_errors, "fallback": False, } # 准备下一轮反馈 if attempt < max_retries: feedback = _build_feedback(combined_score, all_errors, attempt) feedback_history.append(feedback) # 小幅延迟避免 API 限流 if attempt < max_retries: time.sleep(0.5) # 所有重试都失败 → 尝试 best_effort_fallback fallback_data = best_data fallback_raw = best_raw if fallback_data is None and fallback_raw: # 尝试用 best_effort 从最佳原始输出恢复 fallback_data = best_effort_fallback(fallback_raw) # 构建最终结果 result = { "ok": fallback_data is not None, "data": fallback_data, "raw": fallback_raw, "attempts": max_retries, "scores": scores, "errors": ["所有重试均未达到阈值", f"最高分: {best_score}/{threshold}"], "fallback": True, } # 产生告警(通过 stderr 输出,可被 cron 捕获) task_tag = f"[{task_id}]" if task_id else "" print( f"[JSON_VALIDATION] WARNING{task_tag}: " f"{max_retries}次重试后最高分{best_score}/{threshold}," f"使用{'fallback数据' if fallback_data else '空数据'}", file=sys.stderr, flush=True ) # 记录严重错误到 error_lib if error_lib and task_id: error_lib.record( task_id=task_id, attempt=max_retries + 1, score=best_score, errors=["ALL_RETRIES_EXHAUSTED", f"最高分{best_score}/{threshold}"], context=context, is_failure=True, ) return result def _validate_min_values(schema): """确保 schema 包含 min_values 约束""" if schema and "min_values" not in schema: schema["min_values"] = {} return schema def _build_feedback(score, errors, attempt): """构建格式反馈字符串(给 LLM 的下一轮提示)""" parts = [f"--- Attempt {attempt} feedback (score: {score}/100) ---"] if errors: parts.append("Issues to fix:") for e in errors[:5]: # 最多反馈 5 个问题 parts.append(f" - {e}") else: parts.append("No specific errors found, but score below threshold.") return "\n".join(parts) # ═══════════════════════════════════════════════════════════ # Best-effort Fallback # ═══════════════════════════════════════════════════════════ def best_effort_fallback(raw_text): """尽力从非法 JSON 中恢复数据 策略(按优先级): 1. 尝试 auto_fix 后解析 2. 正则提取所有 key:value 对 3. 提取股票代码等关键信息 """ if not raw_text: return None # 策略1: auto_fix fixed, _ = _auto_fix_json(raw_text) if fixed: try: return json.loads(fixed) except json.JSONDecodeError: pass # 策略2: 正则提取所有字段 data = {} # 提取代码 code_match = re.search(r'"code"\s*:\s*"(\d+)"', raw_text) if code_match: data["code"] = code_match.group(1) # 提取名称 name_match = re.search(r'"name"\s*:\s*"([^"]+)"', raw_text) if name_match: data["name"] = name_match.group(1) # 提取数值字段 for field in ["stop_loss", "take_profit", "entry_low", "entry_high", "price", "rr_ratio", "cost"]: pattern = rf'"{field}"\s*:\s*([\d.]+)' m = re.search(pattern, raw_text) if m: try: data[field] = float(m.group(1)) except ValueError: pass # 提取字符串字段 for field in ["action", "timing_signal", "currency", "status"]: pattern = rf'"{field}"\s*:\s*"([^"]*)"' m = re.search(pattern, raw_text) if m: data[field] = m.group(1) return data if data else None # ═══════════════════════════════════════════════════════════ # Format Error Library # ═══════════════════════════════════════════════════════════ class FormatErrorLibrary: """格式错误库 — 持久化记录每次格式错误 每月/周自动汇总,支持按错误类型和任务ID查询。 """ def __init__(self, path=None): self.path = path or ERROR_LIBRARY_PATH self._records = [] self._load() def _load(self): """从磁盘加载已有记录""" if os.path.exists(self.path): try: with open(self.path, "r", encoding="utf-8") as f: data = json.load(f) self._records = data.get("records", []) except (json.JSONDecodeError, Exception): self._records = [] else: self._records = [] def _save(self): """持久化到磁盘""" os.makedirs(os.path.dirname(self.path), exist_ok=True) with open(self.path, "w", encoding="utf-8") as f: json.dump({ "library_version": "1.0", "total_records": len(self._records), "last_updated": datetime.now().isoformat(), "records": self._records[-1000:], # 最多保留最近1000条 }, f, ensure_ascii=False, indent=2) def record(self, task_id, attempt, score, errors, context=None, is_failure=False): """记录一条格式错误""" record = { "task_id": task_id, "timestamp": datetime.now().isoformat(), "attempt": attempt, "score": score, "errors": errors[:10], # 最多保留10个错误 "is_failure": is_failure, } if context: # 只记录上下文的关键信息,避免敏感数据 safe_ctx = {} for k in ["code", "name", "stock"]: if k in (context or {}): safe_ctx[k] = context[k] if safe_ctx: record["context"] = safe_ctx self._records.append(record) self._save() def get_summary(self, days=7): """获取最近 N 天的错误汇总 返回: { "total": 总记录数, "by_type": {"missing_field": N, "syntax_error": N, ...}, "by_task": {"task_id": N}, "worst_tasks": [最差任务列表], } """ cutoff = datetime.now() - timedelta(days=days) recent = [ r for r in self._records if datetime.fromisoformat(r["timestamp"]) >= cutoff ] if not recent: return {"total": 0, "by_type": {}, "by_task": {}, "worst_tasks": []} by_type = defaultdict(int) by_task = defaultdict(int) task_scores = defaultdict(list) for r in recent: task = r.get("task_id", "unknown") by_task[task] += 1 task_scores[task].append(r.get("score", 0)) for err in r.get("errors", []): # 归类错误类型 if "缺少" in err or "missing" in err.lower(): by_type["missing_field"] += 1 elif "syntax" in err.lower() or "JSON" in err or "json" in err: by_type["syntax_error"] += 1 elif "trailing comma" in err.lower() or "尾随逗号" in err: by_type["trailing_comma"] += 1 elif "truncat" in err.lower() or "截断" in err: by_type["truncation"] += 1 else: by_type["other"] += 1 # 最差任务(平均分最低) worst = sorted( [(t, sum(s)/len(s), by_task[t]) for t, s in task_scores.items()], key=lambda x: x[1] )[:5] return { "total": len(recent), "period_days": days, "by_type": dict(by_type), "by_task": dict(by_task), "worst_tasks": [ {"task": t, "avg_score": round(avg, 1), "count": c} for t, avg, c in worst ], "generated_at": datetime.now().isoformat(), } def get_recent(self, limit=20): """获取最近的错误记录""" return list(reversed(self._records[-limit:])) def clear_old(self, keep_days=30): """清理超过 keep_days 的旧记录""" cutoff = datetime.now() - timedelta(days=keep_days) before = len(self._records) self._records = [ r for r in self._records if datetime.fromisoformat(r["timestamp"]) >= cutoff ] after = len(self._records) self._save() return before - after def auto_remediate_prompt(self): """自动分析最近错误模式,返回改进提示词的建议 返回: dict {建议类型: 建议内容} """ summary = self.get_summary(days=7) suggestions = {} by_type = summary.get("by_type", {}) total = summary.get("total", 0) if total == 0: return {"info": "最近7天无格式错误,系统运行良好"} # 字段缺失问题 → 建议在 prompt 中显式列出字段 missing = by_type.get("missing_field", 0) if missing > 0 and missing / max(total, 1) > 0.3: suggestions["prompt_field_emphasis"] = ( f"最近7天 {missing}/{total} 错误是字段缺失," "建议在 prompt 中显式列出所有必需字段的说明和顺序" ) # 语法错误 → 建议使用结构化生成模式 syntax = by_type.get("syntax_error", 0) if syntax > 0 and syntax / max(total, 1) > 0.5: suggestions["structured_generation"] = ( f"最近7天 {syntax}/{total} 错误是 JSON 语法错误," "建议启用 constrained decoding 或 JSON mode" ) # 截断 → 检查 max_tokens trunc = by_type.get("truncation", 0) if trunc > 0: suggestions["increase_max_tokens"] = ( f"最近7天 {trunc} 次截断,建议 max_tokens 从当前值上调" ) if not suggestions: suggestions["info"] = "错误率在可接受范围内,无需自动调整" return suggestions # ═══════════════════════════════════════════════════════════ # 便捷封装: 在 LLM 响应中提取 JSON 的统一入口 # ═══════════════════════════════════════════════════════════ def extract_json_safe(llm_response, schema=None, task_id=None): """安全地从 LLM 响应中提取 JSON,带重试 这是一个单轮快速提取版本(不做重试),适用于不需要 完整 scored_retry_loop 的场景。 返回: (parsed_dict_or_None, score, errors) """ score, errors, parsed = validate_json_syntax(llm_response) if parsed is not None and schema: schema_score, schema_errors = validate_json_schema(parsed, schema) score = int(score * 0.6 + schema_score * 0.4) errors.extend(schema_errors) if parsed is None: # 尝试 fallback parsed = best_effort_fallback(llm_response) return parsed, score, errors # ═══════════════════════════════════════════════════════════ # CLI 入口(测试/调试用) # ═══════════════════════════════════════════════════════════ if __name__ == "__main__": import sys if len(sys.argv) > 1 and sys.argv[1] == "--summary": lib = FormatErrorLibrary() summary = lib.get_summary(days=7) print(json.dumps(summary, ensure_ascii=False, indent=2)) sys.exit(0) if len(sys.argv) > 1 and sys.argv[1] == "--recent": lib = FormatErrorLibrary() recent = lib.get_recent(limit=10) print(json.dumps(recent, ensure_ascii=False, indent=2)) sys.exit(0) if len(sys.argv) > 1 and sys.argv[1] == "--test": # 自测:验证各种情况 test_cases = [ ('{"code": "000001", "name": "平安"}', "标准 JSON"), ('{"code": "000001", "name": "平安",}', "尾随逗号"), ('```json\n{"code": "000001"}\n```', "代码块包裹"), ('{"code": "000001"}', "structured_data 标签"), ('{"code": "000001",\n"name": "平\n安"}', "未转义换行符"), ('{"code": "000001"}}', "多余右括号"), ('{"code": "000001", "price": 12.5, "stop_loss": null, "action": ""}', "空字段"), ] for text, desc in test_cases: score, errors, parsed = validate_json_syntax(text) status = "✅" if parsed else "❌" print(f"{status} {desc}: score={score}") if errors: for e in errors: print(f" - {e}") if parsed: print(f" parsed: {json.dumps(parsed, ensure_ascii=False)[:80]}") print() # 测试 Schema 校验 print("=== Schema 校验测试 ===") test_obj = { "code": "000001", "name": "平安", "action": "持有", "stop_loss": 10.5, "take_profit": 12.0, "entry_low": 10.0, "entry_high": 11.0, "timing_signal": "持有", "rr_ratio": 2.0, } score, errors = validate_json_schema(test_obj) print(f"完整对象: score={score}, errors={errors}") test_obj2 = {"code": "000001", "name": "平安"} score2, errors2 = validate_json_schema(test_obj2) print(f"缺失字段: score={score2}, errors={errors2}") sys.exit(0) # 默认:从 stdin 读取 JSON 文本并验证 text = sys.stdin.read() score, errors, parsed = validate_json_syntax(text) result = { "score": score, "errors": errors, "parsed": parsed, "valid": parsed is not None, } print(json.dumps(result, ensure_ascii=False, indent=2))