黑盒,白盒处理合并

This commit is contained in:
2026-08-24 18:34:17 +08:00
parent 5a968b4e37
commit cbaf3deb36
75 changed files with 10020 additions and 107 deletions
+57
View File
@@ -0,0 +1,57 @@
__version__ = "0.1.0"
from agent.input_parser import InputParser
from agent.rule_loader import RuleLoader
from agent.prompt_builder import PromptBuilder
from agent.api_client import APIClient
from agent.output_writer import OutputWriter
from agent.models import ProgramMeta, FileInfo, CopyField, KeyInfo, TableColumn, TableInfo
def generate(design_md: str, source_cbl: str, file_db_md: str,
cpy_dir: str, db_md: str, output_dir: str = "output",
api_key: str = "sk-6156cccdc9c14d949cf5bfc5afc67a03",
api_model: str = "deepseek-v4-flash",
rules_dir: str = "rules",
max_tokens: int = 32768) -> dict:
"""生成测试数据的主入口函数。"""
import os
print(f"== 解析入力: {design_md}")
parser = InputParser(design_md, source_cbl, file_db_md, cpy_dir, db_md)
meta = parser.run()
if meta.pgm_type == 'サブ':
raise ValueError(f"程序 {meta.program_id} はサブプログラムです。主プログラムのみ処理対象です。")
print(f" プログラムID: {meta.program_id}, パターン: {meta.pgm_pattern}, 入力タイプ: {meta.input_type}")
# If rules_dir is relative, resolve from this file's location or cwd
if not os.path.isabs(rules_dir):
rules_dir = os.path.join(os.path.dirname(__file__), '..', rules_dir)
loader = RuleLoader(rules_dir)
rules_text, group_descriptions, group_count = loader.load(meta)
print(f" ルール読み込み完了, グループ数: {group_count}")
builder = PromptBuilder()
prompt = builder.build(meta, rules_text, group_descriptions, group_count)
client = APIClient(api_key=api_key, model=api_model, max_tokens=max_tokens)
print(f" API呼び出し中...")
result = client.generate(prompt)
print(f" API応答受信")
writer = OutputWriter(output_dir)
output_files = writer.write(meta.program_id, result, meta.input_type)
print(f" 出力完了: {len(output_files)} ファイル")
return {
"output_files": output_files,
"program_id": meta.program_id,
"groups": group_count,
"input_type": meta.input_type,
}
+132
View File
@@ -0,0 +1,132 @@
import json
import time
from typing import Dict, Any, Optional
import requests
class APIClient:
"""DeepSeek API 客户端,含重试逻辑。"""
def __init__(self, api_key: str, model: str = 'deepseek-v4-flash',
base_url: str = 'https://api.deepseek.com/chat/completions',
max_retries: int = 3, timeout: int = 120,
max_tokens: int = 32768):
self.api_key = api_key
self.model = model
self.base_url = base_url
self.max_retries = max_retries
self.timeout = timeout
self.max_tokens = max_tokens
def generate(self, prompt: str) -> Dict[str, Any]:
"""发送 prompt 并返回 AI 生成的结果。"""
system_prompt = (
"你是COBOL程序的测试数据生成专家。"
"请严格按照提供的规则,生成符合格式要求的测试数据。"
"输出必须是可被json.loads()直接解析的JSON,不要包裹在```json```代码块中。"
"不要在JSON前后添加任何说明文字。"
)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt},
]
last_error = None
for attempt in range(1, self.max_retries + 1):
try:
response = self._call_api(messages)
choice = response['choices'][0]
content = choice['message']['content']
if choice.get('finish_reason') == 'length':
# 输出被截断:追加"压缩输出"指示后重试,避免再次浪费全部 max_tokens
last_error = RuntimeError(
"API输出被截断(finish_reason=length),已追加压缩指示重试。"
)
error_msg = (
"前回の出力は長すぎて途中で切れました(finish_reason=length)。"
"レコード数・フィールド数を最小限にし、FILLERなどの冗長な項目は省略して、"
"必ず完結したJSONのみを出力してください。コードブロック(```)で囲まないでください。"
)
messages.append({"role": "assistant", "content": content})
messages.append({"role": "user", "content": error_msg})
continue
data = self._parse_json(content)
if data is not None:
return data
last_error = RuntimeError("API返回内容无法解析为JSON")
error_msg = (
f"前回の出力は有効なJSONではありませんでした。"
f"必ず有効なJSONのみを出力してください。"
f"コードブロック(```)で囲まないでください。"
)
messages.append({"role": "assistant", "content": content})
messages.append({"role": "user", "content": error_msg})
except requests.exceptions.RequestException as e:
last_error = e
if attempt < self.max_retries:
time.sleep(2 ** attempt)
continue
raise RuntimeError(
f"API调用失败,已重试{self.max_retries}次。"
f"最后错误: {last_error}"
)
def _call_api(self, messages: list) -> dict:
"""单次 API 调用。"""
headers = {
'Authorization': f'Bearer {self.api_key}',
'Content-Type': 'application/json',
}
payload = {
'model': self.model,
# deepseek-v4 默认开启思考模式,会把 max_tokens 几乎全部消耗在
# reasoning_content 上,导致 content 为空/截断。测试数据生成是
# 结构化 JSON 输出,关闭思考模式后全部 token 用于正文。
'thinking': {'type': 'disabled'},
'messages': messages,
'temperature': 0.3,
'max_tokens': self.max_tokens,
}
resp = requests.post(
self.base_url,
headers=headers,
json=payload,
timeout=self.timeout,
proxies={"http": None, "https": None},
)
resp.raise_for_status()
return resp.json()
@staticmethod
def _parse_json(text: str) -> Optional[dict]:
"""尝试从文本中提取 JSON。"""
text = text.strip()
if text.startswith('```json'):
text = text[7:]
if text.startswith('```'):
text = text[3:]
if text.endswith('```'):
text = text[:-3]
text = text.strip()
try:
return json.loads(text)
except json.JSONDecodeError:
start = text.find('{')
end = text.rfind('}')
if start >= 0 and end > start:
try:
return json.loads(text[start:end + 1])
except json.JSONDecodeError:
pass
return None
+330
View File
@@ -0,0 +1,330 @@
# agent/input_parser.py
import re
import os
from typing import List, Dict, Optional, Tuple
from agent.models import (
FileInfo, KeyInfo, ModuleInfo, ProgramMeta, CopyField, TableColumn, TableInfo
)
from agent.markdown_utils import (
extract_section, parse_table_rows, parse_table_from_section, find_row_by_key
)
class InputParser:
"""解析 COBOL 程序的详细设计书和相关文件。"""
def __init__(self, design_md_path: str, source_cbl_path: str,
file_db_md_path: str, cpy_dir: str, db_md_path: str):
self.design_md_path = design_md_path
self.source_cbl_path = source_cbl_path
self.file_db_md_path = file_db_md_path
self.cpy_dir = cpy_dir
self.db_md_path = db_md_path
self._design_text = ''
self._source_text = ''
def run(self) -> ProgramMeta:
"""执行完整解析,返回 ProgramMeta。"""
self._design_text = self._read_file(self.design_md_path)
self._source_text = self._read_file(self.source_cbl_path)
meta = ProgramMeta(
program_id='', program_name='', system_name='',
pgm_type='', pgm_pattern='',
summary_lines=[], prerequisites=[],
files=[], keys=[], modules=[],
process_detail='', output_records='',
input_type='file',
copy_fields={}, db_tables={}
)
self._parse_basic_info(meta)
self._parse_use_files(meta)
self._parse_keys(meta)
self._parse_modules(meta)
self._parse_process_detail(meta)
self._parse_output_records(meta)
self._determine_input_type(meta)
self._parse_copybooks(meta)
self._parse_db_definition(meta)
return meta
def _parse_copybooks(self, meta: ProgramMeta):
replacing_list = self._extract_copy_replacing()
for f in meta.files:
if not f.copy_group:
continue
copy_file = os.path.join(self.cpy_dir, f.copy_group + '.cpy')
if not os.path.exists(copy_file):
continue
prefix = f.identifier
for i, (cpy_name, rep_val) in enumerate(replacing_list):
if cpy_name == f.copy_group and rep_val == f.identifier:
prefix = rep_val
replacing_list.pop(i)
break
else:
for i, (cpy_name, rep_val) in enumerate(replacing_list):
if cpy_name == f.copy_group:
prefix = rep_val
replacing_list.pop(i)
break
fields = self._parse_single_copybook(copy_file, prefix)
if f.identifier not in meta.copy_fields:
meta.copy_fields[f.identifier] = fields
def _parse_db_definition(self, meta: ProgramMeta):
"""解析 DB 定义书 .md,提取表结构信息。
只在输入类型涉及 DB 时调用。
"""
if meta.input_type not in ('db', 'mixed'):
return
db_text = self._read_file(self.db_md_path)
sections = re.split(r'\n# ', db_text)
for section in sections:
table_match = re.match(r'^([\w\-]+)', section)
if not table_match:
continue
table_name = table_match.group(1)
section_with_header = '## ' + section
db_info_text = extract_section(section_with_header, 'DB基本情報')
db_info_rows = parse_table_rows(db_info_text)
db_id = ''
copy_id = ''
for row in db_info_rows:
db_id = row.get('DB ID', '') or db_id
copy_id = row.get('COPY ID', '') or copy_id
column_text = extract_section(section_with_header, 'カラム定義')
column_rows = parse_table_rows(column_text)
if not column_rows:
continue
columns = []
pk_columns = []
for row in column_rows:
try:
no = int(row.get('No', '0'))
except ValueError:
continue
is_pk = row.get('PK', '') == ''
nullable = row.get('NULL', 'NULL許可') == 'NULL許可'
col = TableColumn(
no=no,
name_jp=row.get('項目名', ''),
name_en=row.get('項目名(英字名)', ''),
type=row.get('TYPE', ''),
max_len=row.get('最大長', ''),
decimal_digits=row.get('小数桁', ''),
byte_count=row.get('バイト数', ''),
nullable=nullable,
is_pk=is_pk
)
columns.append(col)
if is_pk:
pk_columns.append(col.name_en)
table = TableInfo(
table_name=table_name,
db_id=db_id,
copy_id=copy_id,
columns=columns,
pk_columns=pk_columns
)
meta.db_tables[table_name] = table
def _extract_copy_replacing(self) -> List[Tuple[str, str]]:
"""从 COBOL 源码提取 COPY ... REPLACING ... 映射。
Returns: [(copy_name, replacing_value), ...] 按出现顺序保存。
"""
result = []
pattern = r'COPY\s+(\S+)\s+REPLACING\s+==\(A\)==\s+BY\s+==(\S+)=='
for match in re.finditer(pattern, self._source_text, re.IGNORECASE):
result.append((match.group(1), match.group(2)))
return result
def _parse_single_copybook(self, copy_path: str, prefix: str) -> List[CopyField]:
copy_text = self._read_file(copy_path)
fields = []
for line in copy_text.split('\n'):
line = line.strip()
if not line or line.startswith('*'):
continue
match = re.match(
r'(\d{2})\s+\(A\)(\S+)\s+PIC\s+(.+?)\.',
line, re.IGNORECASE
)
if not match:
continue
level = int(match.group(1))
raw_name = '(A)' + match.group(2)
pic_full = match.group(3).strip()
actual_name = prefix + '-' + match.group(2)
pic_bytes = self._calculate_pic_bytes(pic_full)
fields.append(CopyField(
level=level,
name=actual_name,
raw_name=raw_name,
pic_type=pic_full,
pic_bytes=pic_bytes
))
return fields
@staticmethod
def _calculate_pic_bytes(pic: str) -> int:
pic_upper = pic.upper().strip()
if 'COMP-3' in pic_upper:
base = pic_upper.replace('COMP-3', '').strip()
total_digits = sum(int(num) for num in re.findall(r'9\((\d+)\)', base))
if base.startswith('S'):
return (total_digits + 2) // 2
else:
return (total_digits + 1) // 2
if 'COMP' in pic_upper or 'BINARY' in pic_upper:
base = re.sub(r'\s*(COMP|BINARY)\s*', '', pic_upper).strip()
total_digits = sum(int(num) for num in re.findall(r'9\((\d+)\)', base))
if total_digits <= 4:
return 2
elif total_digits <= 9:
return 4
else:
return 8
total = 0
for num in re.findall(r'X\((\d+)\)', pic_upper):
total += int(num)
for num in re.findall(r'G\((\d+)\)', pic_upper):
total += int(num) * 2
for num in re.findall(r'9\((\d+)\)', pic_upper):
total += int(num)
if total > 0:
return total
return 0
@staticmethod
def _read_file(path: str) -> str:
with open(path, 'r', encoding='utf-8') as f:
return f.read()
def _parse_basic_info(self, meta: ProgramMeta):
rows = parse_table_from_section(self._design_text, '基本情報')
for row in rows:
item = row.get('項目', '')
value = row.get('内容', '')
if item == 'システム名':
meta.system_name = value
elif item == 'プログラムID':
meta.program_id = value
elif item == 'プログラム名':
meta.program_name = value
elif item == 'PGMタイプ':
meta.pgm_type = value
elif item == 'PGMパターン':
meta.pgm_pattern = value
elif item == '機能概要':
meta.summary_lines.append(value)
elif item == '':
if value:
meta.summary_lines.append(value)
def _parse_use_files(self, meta: ProgramMeta):
rows = parse_table_from_section(self._design_text, '使用ファイル一覧')
for row in rows:
try:
no_str = row.get('NO', '0')
try:
no = int(no_str) if no_str and no_str.strip() not in ('', '') else 0
except ValueError:
no = 0
rec_len_str = row.get('レコード長', '0')
try:
rec_len = int(rec_len_str) if rec_len_str and rec_len_str.strip() not in ('', '') else 0
except ValueError:
rec_len = 0
f = FileInfo(
no=no,
file_db_name=row.get('使用ファイル/DB名', ''),
identifier=row.get('識別子', ''),
dd_name=row.get('DD名', ''),
io=row.get('I/O', ''),
copy_group=row.get('COPY群', ''),
format=row.get('形式', ''),
record_len=rec_len,
medium=row.get('媒体', ''),
remarks=row.get('備考', '')
)
meta.files.append(f)
except (ValueError, KeyError):
continue
def _parse_keys(self, meta: ProgramMeta):
rows = parse_table_from_section(self._design_text, 'キー項目一覧')
for row in rows:
try:
k = KeyInfo(
no=int(row.get('NO', '0')),
file_name=row.get('ファイル名', ''),
sort_condition=row.get('ソート条件(キー項目)', ''),
key_condition=row.get('キー条件(マッチング/キーブレイク)', '')
)
meta.keys.append(k)
except (ValueError, KeyError):
continue
def _parse_modules(self, meta: ProgramMeta):
rows = parse_table_from_section(self._design_text, '使用モジュール一覧')
for row in rows:
try:
m = ModuleInfo(
no=int(row.get('NO', '0')),
function=row.get('機能', ''),
program_id=row.get('プログラムID', ''),
copy_name=row.get('使用COPY名', '')
)
meta.modules.append(m)
except (ValueError, KeyError):
continue
def _parse_process_detail(self, meta: ProgramMeta):
meta.process_detail = extract_section(self._design_text, '処理詳細')
def _parse_output_records(self, meta: ProgramMeta):
meta.output_records = extract_section(self._design_text, '出力レコード定義')
def _determine_input_type(self, meta: ProgramMeta):
input_mediums = set()
for f in meta.files:
if 'I' in f.io:
input_mediums.add(f.medium)
if not input_mediums:
meta.input_type = 'file'
elif input_mediums == {'PS'}:
meta.input_type = 'file'
elif input_mediums == {'DB'}:
meta.input_type = 'db'
else:
meta.input_type = 'mixed'
@@ -0,0 +1,63 @@
# agent/markdown_utils.py
import re
from typing import List, Dict, Optional
def extract_section(md_text: str, section_title: str) -> str:
"""Extract content of a ## or ### section from markdown text.
Stops at the next heading of equal or higher level (fewer or equal #'s).
For ## sections, includes ### sub-sections.
"""
heading_re = rf'^(#{{2,3}})\s+{re.escape(section_title)}[^\S\n]*$'
h_match = re.search(heading_re, md_text, re.MULTILINE)
if not h_match:
return ''
level = len(h_match.group(1))
start = h_match.end()
next_re = rf'^#{{1,{level}}}\s'
n_match = re.search(next_re, md_text[start:], re.MULTILINE)
end = start + n_match.start() if n_match else len(md_text)
return md_text[start:end]
def parse_table_rows(text: str) -> List[Dict[str, str]]:
"""Parse a markdown table from text, return list of row dicts.
Handles tables with exactly one header row and one separator row.
"""
lines = []
for line in text.split('\n'):
stripped = line.strip()
if stripped.startswith('|') and stripped.endswith('|'):
if re.match(r'^\|[\s\-:]+\|', stripped):
continue
lines.append(stripped)
if len(lines) < 1:
return []
headers = [cell.strip() for cell in lines[0].split('|')[1:-1]]
rows = []
for line in lines[1:]:
cells = [cell.strip() for cell in line.split('|')[1:-1]]
if len(cells) == len(headers):
rows.append(dict(zip(headers, cells)))
return rows
def parse_table_from_section(md_text: str, section_title: str) -> List[Dict[str, str]]:
"""Find a ### section and parse its first table."""
section_text = extract_section(md_text, section_title)
if not section_text:
return []
return parse_table_rows(section_text)
def find_row_by_key(rows: List[Dict[str, str]], key_col: str, key_value: str) -> Optional[Dict[str, str]]:
"""Find a table row where a specific column matches the key value."""
for row in rows:
if row.get(key_col, '').strip() == key_value:
return row
return None
+89
View File
@@ -0,0 +1,89 @@
from dataclasses import dataclass, field
from typing import List, Dict, Optional
@dataclass
class FileInfo:
"""使用ファイル一覧 中的一行"""
no: int
file_db_name: str
identifier: str
dd_name: str
io: str
copy_group: str
format: str
record_len: int
medium: str
remarks: str
@dataclass
class CopyField:
"""COPYBOOK 中的单个字段"""
level: int
name: str
raw_name: str
pic_type: str
pic_bytes: int
@dataclass
class KeyInfo:
"""キー項目一覧 中的一行"""
no: int
file_name: str
sort_condition: str
key_condition: str
@dataclass
class ModuleInfo:
"""使用モジュール一覧 中的一行"""
no: int
function: str
program_id: str
copy_name: str
@dataclass
class TableColumn:
"""DB 表中的单个字段"""
no: int
name_jp: str
name_en: str
type: str
max_len: str
decimal_digits: str
byte_count: str
nullable: bool
is_pk: bool
@dataclass
class TableInfo:
"""DB 表定义"""
table_name: str
db_id: str
copy_id: str
columns: List[TableColumn]
pk_columns: List[str]
@dataclass
class ProgramMeta:
"""程序完整元数据"""
program_id: str
program_name: str
system_name: str
pgm_type: str
pgm_pattern: str
summary_lines: List[str]
prerequisites: List[Dict[str, str]]
files: List[FileInfo]
keys: List[KeyInfo]
modules: List[ModuleInfo]
process_detail: str
output_records: str
input_type: str
copy_fields: Dict[str, List[CopyField]]
db_tables: Dict[str, TableInfo]
@@ -0,0 +1,75 @@
import json
import os
from typing import Dict, Any
class OutputWriter:
"""将 AI 生成的数据写入文件系统。"""
def __init__(self, output_dir: str):
self.output_dir = output_dir
def write(self, program_id: str, ai_result: Dict[str, Any],
input_type: str) -> Dict[str, str]:
"""写入所有组的输出文件。
Returns:
{group_folder: written_file_path} 映射
"""
written = {}
groups = ai_result.get('groups', ai_result)
for group_key in sorted(groups.keys()):
group_data = groups[group_key]
group_dir = os.path.join(self.output_dir, program_id, 'black_box', group_key)
os.makedirs(group_dir, exist_ok=True)
if input_type in ('file', 'mixed'):
json_path = self._write_json(group_dir, program_id, group_key, group_data)
written[f"{group_key}/json"] = json_path
if input_type in ('db', 'mixed'):
sql_path = self._write_sql(group_dir, program_id, group_key, group_data)
written[f"{group_key}/sql"] = sql_path
summary_md = ai_result.get('summary_md')
if summary_md:
md_path = self._write_case_table(program_id, summary_md)
written["case表"] = md_path
return written
def _write_case_table(self, program_id: str, content: str) -> str:
"""写入 case表.md(组与规则md的对应关系说明)。"""
black_box_dir = os.path.join(self.output_dir, program_id, 'black_box')
os.makedirs(black_box_dir, exist_ok=True)
filepath = os.path.join(black_box_dir, 'case表.md')
with open(filepath, 'w', encoding='utf-8') as f:
f.write(content)
return filepath
def _write_json(self, group_dir: str, program_id: str,
group_key: str, data: Any) -> str:
"""写入 JSON 文件。"""
filename = f"{program_id}_{group_key}.json"
filepath = os.path.join(group_dir, filename)
with open(filepath, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=2)
return filepath
def _write_sql(self, group_dir: str, program_id: str,
group_key: str, data: Any) -> str:
"""写入 SQL 文件。"""
filename = f"{program_id}_{group_key}.sql"
filepath = os.path.join(group_dir, filename)
sql_content = data if isinstance(data, str) else data.get('sql', json.dumps(data, ensure_ascii=False))
with open(filepath, 'w', encoding='utf-8') as f:
f.write(sql_content)
return filepath
@@ -0,0 +1,155 @@
from typing import List
from agent.models import ProgramMeta
class PromptBuilder:
"""将解析后的程序元数据和规则组装成 API prompt。"""
def build(self, meta: ProgramMeta, rules_text: str,
group_descriptions: List[str], group_count: int) -> str:
"""构建完整的 API prompt。"""
parts = []
parts.append(self._build_basic_info(meta))
parts.append(self._build_process_detail(meta))
parts.append(self._build_input_structures(meta))
parts.append(self._build_output_records(meta))
parts.append(rules_text)
parts.append(self._build_output_format(meta))
parts.append(self._build_generation_instruction(
group_descriptions, group_count, meta.input_type
))
return '\n\n'.join(parts)
def _build_basic_info(self, meta: ProgramMeta) -> str:
lines = [
f"## プログラム基本情報",
f"- システム名: {meta.system_name}",
f"- プログラムID: {meta.program_id}",
f"- プログラム名: {meta.program_name}",
f"- PGMパターン: {meta.pgm_pattern}",
f"- 入力タイプ: {self._input_type_label(meta.input_type)}",
]
if meta.summary_lines:
lines.append(f"- 機能概要: {' '.join(meta.summary_lines)}")
return '\n'.join(lines)
def _build_process_detail(self, meta: ProgramMeta) -> str:
return f"## 処理詳細\n\n```\n{meta.process_detail}\n```"
def _build_input_structures(self, meta: ProgramMeta) -> str:
parts = ["## 入力構造"]
input_files = [f for f in meta.files if 'I' in f.io]
for f in input_files:
parts.append(f"### ファイル {f.identifier} (DD名: {f.dd_name}, COPY: {f.copy_group}, 媒体: {f.medium})")
fields = meta.copy_fields.get(f.identifier, [])
if fields:
parts.append("| 項目名 | PIC | バイト数 |")
parts.append("|--------|-----|----------|")
for cf in fields:
parts.append(f"| {cf.name} | {cf.pic_type} | {cf.pic_bytes} |")
else:
parts.append("(構造情報なし)")
db_inputs = [f for f in meta.files if 'I' in f.io and f.medium == 'DB']
if db_inputs or meta.db_tables:
parts.append("### DBテーブル構造")
for table_name, table in meta.db_tables.items():
parts.append(f"**表名: {table_name}**")
pk_str = ', '.join(table.pk_columns)
parts.append(f"主キー: {pk_str}")
parts.append("| 項目名 | 英字名 | タイプ | 最大長 | KEY |")
parts.append("|--------|--------|--------|--------|-----|")
for col in table.columns:
key_mark = '' if col.is_pk else ''
parts.append(f"| {col.name_jp} | {col.name_en} | {col.type} | {col.max_len} | {key_mark} |")
return '\n\n'.join(parts)
def _build_output_records(self, meta: ProgramMeta) -> str:
return f"## 出力レコード定義\n\n```\n{meta.output_records}\n```"
def _build_output_format(self, meta: ProgramMeta) -> str:
lines = [
"## 出力形式",
"",
"### JSON形式(ファイル入力の場合)",
"",
"JSON構造:",
"```json",
"{",
' "program": "{プログラムID}",',
' "records": [',
' { "input": { "FD名": { "項目名": "", ... } } },',
' ...',
' ]',
"}",
"```",
"",
"### 項目値のルール(PIC → JSON値)",
"| PIC | JSON内表示 | 例 |",
"|-----|-----------|-----|",
"| PIC X(n) | 左詰め + スペース埋め | `\"A0000001\"` |",
"| PIC 9(n) | 右詰め + 先行ゼロ | `\"00000101\"` |",
"| PIC S9(n) | 符号 + 右詰め + 先行ゼロ | `\"+0000101\"` |",
"| PIC S9(n)V9(m) | 符号 + 右詰め + 小数点含む | `\"+001234567\"` |",
"| PIC S9(n) COMP | 通常の10進数文字列 | `\"300\"` |",
"| PIC S9(n) COMP-3 | 通常の10進数文字列 | `\"1234\"` |",
"| PIC 9(n) COMP-3 | 通常の10進数文字列 | `\"1234\"` |",
"| FILLER(純粋予約) | グループIDとレコード番号を含むパターン | `\"D000000...001\"` |",
"| FILLER(業務予約) | 全スペース(PIC X) または 全ゼロ(PIC 9) | |",
"",
"### SQL形式(DB入力の場合)",
"各グループ1つのSQLファイル: `{program}_g{groupId}.sql`",
"",
"SQL例:",
"```sql",
"-- Group: 1",
"INSERT INTO TABLE_NAME (COL1, COL2) VALUES",
"('val1', 'val2');",
"```",
"",
"### データ生成の注意",
"- 項目名に意味がある場合(DATE→日付、NAME→氏名)、実際の形式に合った値を生成すること",
"- 隣接するレコード間で、同じ項目に異なる値を設定すること",
]
return '\n'.join(lines)
def _build_generation_instruction(self, descriptions: List[str],
count: int, input_type: str) -> str:
lines = [
"## 生成指示",
f"生成するグループ数: {count}",
"各グループの内容:",
]
for i, desc in enumerate(descriptions, 1):
lines.append(f" - g{i}: {desc}")
if input_type == 'file':
lines.append("出力: JSONファイル(1グループ=1JSONファイル)")
elif input_type == 'db':
lines.append("出力: SQL INSERTファイル(1グループ=1SQLファイル)")
else:
lines.append("出力: JSONファイル + SQL INSERTファイル")
lines.append("")
lines.append("応答形式: JSONで返してください。最上位キー \"groups\" を使い、各グループを \"g1\", \"g2\" ... とします。")
lines.append("例: {\"groups\": {\"g1\": {\"program\": \"XXX\", \"records\": [...]}, \"g2\": {...}}}")
lines.append("")
lines.append("### case表.md 生成 (summary_md)")
lines.append("最上位に \"summary_md\" フィールドも追加してください。値は Markdown 文字列で、")
lines.append("各グループと規則mdの対応関係を説明する case表 です。")
lines.append("case表 には規則mdに定義された全グループ(g1〜gn)を列挙し、以下の列を含むMarkdownテーブルにしてください:")
lines.append("| 组 | 规则md中的定义/用途 | 本次是否生成 |")
lines.append("- 规则md中定义且本次已生成数据的组:状态列标注 已生成")
lines.append("- 规则md中定义、但本次判断没必要生成而省略的组:状态列标注 ✗")
lines.append("summary_md の最後にはルールファイル名(例: マッチング(1-1).md)を記載してください。")
return '\n'.join(lines)
@staticmethod
def _input_type_label(t: str) -> str:
labels = {'file': 'ファイル', 'db': 'DB', 'mixed': '混合(ファイル+DB'}
return labels.get(t, t)
+229
View File
@@ -0,0 +1,229 @@
import os
import re
from typing import List, Tuple, Optional
from agent.models import ProgramMeta
PGM_PATTERN_MAP = {
'マッチング(1:1)': 'マッチング(1-1).md',
'マッチング(11': 'マッチング(1-1).md',
'マッチング(1:N)': 'マッチング(1-N).md',
'マッチング(1N': 'マッチング(1-N).md',
'マッチング(1:N) + EVALUATE 4パターン分岐': 'マッチング(1-N).md',
'マッチング(N:1)': 'マッチング(N-1).md',
'マッチング(N1': 'マッチング(N-1).md',
'マッチング(M:N)': 'マッチング(M-N).md',
'マッチングMN⇒出力M件': 'マッチング(M-N).md',
'マッチングMN⇒出力N件': 'マッチング(M-N).md',
'マッチングMN⇒出力M×N件': 'マッチング(M-N).md',
'18M:N→M件マッチング)': 'マッチング(M-N).md',
'19M:N→N件マッチング+集計)': 'マッチング(M-N).md',
'20M:N→M×N件直積出力)': 'マッチング(M-N).md',
'レイアウト編集のみ(GETPUT)': 'レイアウト編集のみ(GETPUT).md',
'レイアウト編集のみ(GETPUT': 'レイアウト編集のみ(GETPUT).md',
'レイアウト編集のみ': 'レイアウト編集のみ(GETPUT).md',
'GETPUT(編集出力)': 'GETPUT(編集出力).md',
'項目チェック': '項目チェック(重複含まず).md',
'項目チェック(重複なし)': '項目チェック(重複含まず).md',
'項目チェック(重複なし)': '項目チェック(重複含まず).md',
'項目チェック(重複含まず)': '項目チェック(重複含まず).md',
'項目チェック(重複含む)': '項目チェック(重複含まず).md',
'項目チェック(半角20桁/4桁)': '項目チェック(半角20桁-4桁).md',
'27(半角20桁/4桁チェック)': '項目チェック(半角20桁-4桁).md',
'半角20桁/4桁': '項目チェック(半角20桁-4桁).md',
'振り分け': '振り分け(IF).md',
'振り分け(IF文)': '振り分け(IF).md',
'振り分け(IF文、EVALUATE文)': '振り分け(IF).md',
'振り分け(EVALUATE文)': '振り分け(EVALUATE).md',
'振り分け(EVALUATE文)': '振り分け(EVALUATE).md',
'キーブレイク': 'キーブレイク(集計).md',
'キーブレイク(集計)': 'キーブレイク(集計).md',
'キーブレイク(集計)': 'キーブレイク(集計).md',
'キーブレイク(集約)': 'キーブレイク(集約).md',
'キーブレイク(集約)': 'キーブレイク(集約).md',
'キーブレイク(集計、集約)': 'キーブレイク(集計).md',
'キーブレイク(集計、集約の以外)': 'キーブレイク(非集計).md',
'キーブレイク(集計集約以外)': 'キーブレイク(非集計).md',
'キーブレイク(集計集約以外)': 'キーブレイク(非集計).md',
'キーブレイク(非集計)': 'キーブレイク(非集計).md',
'1:Nキーブレイク(同キー集約)': 'キーブレイク(集計).md',
'1:N+キーブレイク(同キー)': 'キーブレイク(集計).md',
'1:N+同キーキーブレイク': 'キーブレイク(集計).md',
'321:N+同キーキーブレイク)': 'キーブレイク(集計).md',
'1:N+同キー': 'キーブレイク(集計).md',
'1:N+キーブレイク(異キー)': '1-N+キーブレイク(異キー).md',
'1:N+キーブレイク(異キー)': '1-N+キーブレイク(異キー).md',
'1:N+異キーキーブレイク': '1-N+キーブレイク(異キー).md',
'331:N+異キーキーブレイク)': '1-N+キーブレイク(異キー).md',
'1:N+異キー': '1-N+キーブレイク(異キー).md',
'DB更新': 'DB更新.md',
'DB更新 + SYSIN読込(P28)': 'DB更新.md',
'SELECT処理': 'SELECT処理.md',
'SELECT条件': 'SELECT処理.md',
'SELECT条件': 'SELECT処理.md',
'DB検索': 'SELECT処理.md',
'DB検索': 'SELECT処理.md',
'50分割': '50分割.md',
'50分割': '50分割.md',
'25分割': '25分割.md',
'25分割': '25分割.md',
'1125分割)': '25分割.md',
'100分割': '100分割.md',
'100分割': '100分割.md',
'12100分割)': '100分割.md',
'MERGE(複数ファイル結合)': 'MERGE.md',
'MERGE(複数ファイル結合)': 'MERGE.md',
'CSV→FB変換(改行あり)': 'CSV→FB変換.md',
'CSV→FB変換(改行あり)': 'CSV→FB変換.md',
'CSV→FB変換(改行なし)': 'CSV→FB変換.md',
'CSV→FB変換(改行なし)': 'CSV→FB変換.md',
'2段階マッチング(1:1⇒1:1': '2段階マッチング(1:1⇒1:1).md',
'2段階マッチング(11⇒11': '2段階マッチング(1:1⇒1:1).md',
'2段階マッチング(N:1⇒N:1': '2段階マッチング(N:1⇒N:1).md',
'2段階マッチング(N1⇒N1': '2段階マッチング(N:1⇒N:1).md',
'2段階マッチング(M:N⇒M:N': '2段階マッチング(M:N⇒M:N).md',
'2段階マッチング(MN⇒MN': '2段階マッチング(M:N⇒M:N).md',
'マッチング(1:1→1:1 2段階)': '2段階マッチング(1:1⇒1:1).md',
'マッチング(N:1→N:1 2段階)': '2段階マッチング(N:1⇒N:1).md',
'マッチング(M:N→M:N 2段階)': '2段階マッチング(M:N⇒M:N).md',
'内部テーブル検索': '内部テーブル検索.md',
'24(内部テーブル検索)': '内部テーブル検索.md',
'サブプログラム使用': 'サブプログラム.md',
'サブプログラム': 'サブプログラム.md',
'サブPGM': 'サブプログラム.md',
'25(サブPGM': 'サブプログラム.md',
'SYSIN読込': 'SYSIN読込.md',
'SYSIN読込': 'SYSIN読込.md',
'28SYSIN読込)': 'SYSIN読込.md',
'ASCII→EBCDIC変換': 'ASCII→EBCDIC変換.md',
'ASCII→EBCDIC変換': 'ASCII→EBCDIC変換.md',
'29ASCII→EBCDIC変換)': 'ASCII→EBCDIC変換.md',
'SORTINPUT/OUTPUT PROCEDURE': 'SORT.md',
'34SORT INPUT/OUTPUT PROCEDURE': 'SORT.md',
'SORT': 'SORT.md',
'オンラインPGM': 'オンラインPGM.md',
'オンラインPGM': 'オンラインPGM.md',
'オンライン': 'オンラインPGM.md',
'編集出力(ランキング生成)': 'ランキング生成.md',
'36(ランキング生成)': 'ランキング生成.md',
'ランキング生成': 'ランキング生成.md',
'編集出力(印刷制御)': 'GETPUT(編集出力).md',
'編集出力': 'GETPUT(編集出力).md',
}
SPECIAL_FEATURE_CHECKS = [
(['場合', 'EVALUATE', 'IF'], '条件分支.md'),
]
class RuleLoader:
"""根据程序特征加载对应的数据生成规则。"""
def __init__(self, rules_dir: str):
self.pgm_pattern_dir = os.path.join(rules_dir, 'pgm_pattern')
self.special_feature_dir = os.path.join(rules_dir, 'special_feature')
self.last_rule_file = None
def load(self, meta: ProgramMeta) -> Tuple[str, List[str], int]:
"""加载所有相关规则,返回 (合并后的规则文本, 组描述列表, 组数)。"""
pgm_rule = self._load_pgm_pattern_rule(meta.pgm_pattern)
if pgm_rule is None:
raise FileNotFoundError(
f"PGM模式 '{meta.pgm_pattern}' の規則ファイルが見つかりません。"
f"{self.pgm_pattern_dir} に .md ファイルを追加してください。"
)
group_descriptions, group_count = self._parse_group_info(pgm_rule)
parts = [pgm_rule]
for keywords, rule_file in SPECIAL_FEATURE_CHECKS:
if self._detect_feature(meta.process_detail, keywords):
feature_rule = self._read_rule_file(
os.path.join(self.special_feature_dir, rule_file)
)
if feature_rule:
parts.append(feature_rule)
combined = '\n\n---\n\n'.join(parts)
return combined, group_descriptions, group_count
def _load_pgm_pattern_rule(self, pgm_pattern: str) -> Optional[str]:
"""根据 PGMパターン 加载对应的规则文件。"""
filename = PGM_PATTERN_MAP.get(pgm_pattern)
if filename:
path = os.path.join(self.pgm_pattern_dir, filename)
if os.path.exists(path):
self.last_rule_file = filename
return self._read_rule_file(path)
if os.path.isdir(self.pgm_pattern_dir):
available = sorted(os.listdir(self.pgm_pattern_dir))
# 对每个规则文件名(去掉.md),检查它是否出现在 PGMパターン 中
# 半角/全角括弧に加え、'1:N' vs '1-N' のコロン/ハイフン差異も吸収する
pgm_lower = pgm_pattern.lower().replace(':', '-')
for fname in available:
if fname.endswith('.md'):
rule_name = fname[:-3].lower()
if rule_name in pgm_lower:
self.last_rule_file = fname
return self._read_rule_file(os.path.join(self.pgm_pattern_dir, fname))
return None
@staticmethod
def _read_rule_file(path: str) -> str:
with open(path, 'r', encoding='utf-8') as f:
return f.read()
@staticmethod
def _parse_group_info(rule_text: str) -> Tuple[List[str], int]:
"""从规则文本中解析组信息。
只统计表头首列为""的表格(生成的数据组),
避免把记录明细表中的数字(如 record序号 / 第1条)误算成组号。
"""
descriptions = []
group_count = 0
in_group_table = False
for line in rule_text.split('\n'):
stripped = line.strip()
if not stripped.startswith('|'):
in_group_table = False
continue
# 分隔行(|---|)跳过
if re.match(r'^\|[\s\-:]+\|', stripped):
continue
cells = [c.strip() for c in stripped.split('|')[1:-1]]
if not cells:
continue
# 表头行:首列为"组"时进入组表格
if cells[0] == '':
in_group_table = True
continue
if not in_group_table:
continue
if len(cells) >= 2:
try:
group_num = int(cells[0])
if group_num > group_count:
group_count = group_num
descriptions.append(cells[1])
except ValueError:
pass
return descriptions, group_count
@staticmethod
def _detect_feature(process_detail: str, keywords: List[str]) -> bool:
"""检测処理詳細中是否包含特定关键词。"""
for kw in keywords:
if kw in process_detail:
return True
return False