Files
cobol-java-v3/cobol_testgen/file_io.py
T

276 lines
9.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""COBOL 文件 I/ODISPLAY/COMP/COMP-3 pack/unpack + 文件读写"""
import struct
import logging
logger = logging.getLogger(__name__)
# ── 存储长度 ──
def get_storage_length(field: dict) -> int:
"""返回字段在文件中的字节长度"""
pi = field.get('pic_info', {})
digits = pi.get('digits', 0)
usage = field.get('usage')
if not usage or usage == 'DISPLAY':
l = pi.get('length')
if l:
return l
return digits + pi.get('decimal', 0) or 1
elif usage in ('COMP', 'BINARY'):
if digits <= 2:
return 1
elif digits <= 4:
return 2
elif digits <= 9:
return 4
else:
return 8
elif usage in ('COMP-3', 'PACKED-DECIMAL'):
# 打包小数含小数位(如 9(3)V9(2)=5 位+符号 → 3 字节)。
# 只算整数位(digits)会少算 1 字节,导致与 GnuCOBOL FD 布局错位。
decimal = pi.get('decimal', 0)
return (digits + decimal + 2) // 2
else:
raise ValueError(f"Unsupported USAGE: {usage}")
# ── pack / unpack ──
def _default_value(field: dict) -> str:
"""字段值为空/缺失时的默认值"""
pi = field.get('pic_info', {})
usage = field.get('usage')
if not usage or usage == 'DISPLAY':
total = pi.get('length') or (pi.get('digits', 0) + pi.get('decimal', 0)) or 1
return ' ' * total
digits = pi.get('digits', 0)
return '0' * digits
def pack_value(value: str, field: dict) -> bytes:
"""将 JSON 字符串值编码为二进制文件表示"""
if not value or value.strip() == '':
value = _default_value(field)
usage = field.get('usage')
pi = field.get('pic_info', {})
ptype = pi.get('type', 'unknown')
digits = pi.get('digits', 0)
signed = pi.get('signed', False)
if not usage or usage == 'DISPLAY':
total = pi.get('length') or (digits + pi.get('decimal', 0)) or 1
if ptype in ('numeric', 'numeric-edited'):
s = str(value).zfill(total)
else:
s = str(value).ljust(total)
return s.encode('utf-8')[:total]
int_val = int(str(value).strip())
if usage in ('COMP', 'BINARY'):
size = get_storage_length(field)
fmt_map = {1: 'b', 2: 'h', 4: 'i', 8: 'q'}
fmt = fmt_map[size]
if not signed:
fmt = fmt.upper()
# 本 GnuCOBOLGC32-BDB-SP1、主机兼容)的 COMP 按大端存储。
return struct.pack('>' + fmt, int_val)
elif usage in ('COMP-3', 'PACKED-DECIMAL'):
# 打包长度含小数位(digits + decimal
abs_str = str(abs(int_val)).zfill(digits + pi.get('decimal', 0))
nibbles = [int(ch) for ch in abs_str]
if not signed:
nibbles.append(0xF)
elif int_val >= 0:
nibbles.append(0xC)
else:
nibbles.append(0xD)
if len(nibbles) % 2 == 1:
nibbles.insert(0, 0)
buf = bytearray()
for i in range(0, len(nibbles), 2):
buf.append((nibbles[i] << 4) | nibbles[i + 1])
return bytes(buf)
else:
raise ValueError(f"Unsupported USAGE: {usage}")
def unpack_value(data: bytes, field: dict) -> str:
"""将二进制数据解码为 JSON 字符串值"""
usage = field.get('usage')
pi = field.get('pic_info', {})
digits = pi.get('digits', 0)
signed = pi.get('signed', False)
if not usage or usage == 'DISPLAY':
return data.decode('utf-8').rstrip()
elif usage in ('COMP', 'BINARY'):
size = len(data)
fmt_map = {1: 'b', 2: 'h', 4: 'i', 8: 'q'}
fmt = fmt_map[size]
if not signed:
fmt = fmt.upper()
# 与 pack_value 一致:本 GnuCOBOL 的 COMP 为大端存储。
val = struct.unpack('>' + fmt, data)[0]
sign = '-' if val < 0 else ''
return f"{sign}{str(abs(val)).zfill(digits)}"
elif usage in ('COMP-3', 'PACKED-DECIMAL'):
nibbles = []
for byte in data:
nibbles.append((byte >> 4) & 0x0F)
nibbles.append(byte & 0x0F)
sign = nibbles[-1]
nibbles = nibbles[:-1]
chars = [str(n) for n in nibbles]
num_str = ''.join(chars).lstrip('0') or '0'
if signed and sign == 0xD:
num_str = '-' + num_str
# 补零长度含小数位
return num_str.zfill(digits + pi.get('decimal', 0))
else:
raise ValueError(f"Unsupported USAGE: {usage}")
# ── 文件读写 ──
def compute_record_size(fd_field_dicts: list[dict]) -> int:
"""计算 FD 记录的总字节长度"""
return sum(get_storage_length(f) for f in fd_field_dicts)
def has_any_binary(fd_field_dicts: list[dict]) -> bool:
"""FD 中是否有 COMP/COMP-3 字段"""
for f in fd_field_dicts:
usage = f.get('usage')
if usage and usage not in (None, 'DISPLAY'):
return True
return False
def write_input_file(records: list[dict], fd_field_dicts: list[dict],
output_path: str, line_sequential: bool = False):
"""将记录列表写入 COBOL 输入文件"""
with open(output_path, 'wb') as f:
for record in records:
for field_dict in fd_field_dicts:
val = record.get(field_dict['name'], '')
packed = pack_value(val, field_dict)
f.write(packed)
if line_sequential:
f.write(b'\n')
logger.info(f" wrote {len(records)} records to {output_path}")
def read_output_file(file_path: str, fd_field_dicts: list[dict],
line_sequential: bool = False, recording_mode: str = 'F') -> list[dict]:
"""从 COBOL 输出文件读取记录"""
if recording_mode == 'V':
return _read_variable_file(file_path, fd_field_dicts)
record_size = compute_record_size(fd_field_dicts)
records = []
if line_sequential:
with open(file_path, 'rb') as f:
for raw_line in f:
raw_line = raw_line.rstrip(b'\r\n')
records.append(_unpack_record(raw_line, fd_field_dicts))
else:
record_size = compute_record_size(fd_field_dicts)
with open(file_path, 'rb') as f:
while True:
data = f.read(record_size)
if not data:
break
records.append(_unpack_record(data, fd_field_dicts))
return records
def _read_variable_file(file_path: str, fd_field_dicts: list[dict]) -> list[dict]:
"""读取 RECORDING MODE V 文件。
GnuCOBOL on Linux 可能写入 RDW 前缀,也可能不写入。
先尝试 RDW 方式;如果第一笔的 rec_len 不合理(> 10000),
则降级为固定长度读取。
"""
record_size = compute_record_size(fd_field_dicts)
if record_size == 0:
return []
raw = open(file_path, 'rb').read()
if len(raw) < 4:
return []
first_rdw = int.from_bytes(raw[:2], 'little')
if first_rdw > 10000 or (first_rdw - 4) > record_size * 2:
# 没有 RDW 前缀 → 固定长度读取
records = []
offset = 0
while offset + record_size <= len(raw):
records.append(_unpack_record(raw[offset:offset + record_size], fd_field_dicts))
offset += record_size
return records
# 正常 RDW 方式
records = []
offset = 0
while offset < len(raw):
if offset + 4 > len(raw):
break
rdw_len = int.from_bytes(raw[offset:offset + 2], 'little')
data_len = rdw_len - 4 if rdw_len >= 4 else 0
offset += 4
if offset + data_len > len(raw):
break
records.append(_unpack_record(raw[offset:offset + data_len], fd_field_dicts))
offset += data_len
return records
def _unpack_record(data: bytes, fd_field_dicts: list[dict]) -> dict:
"""从字节数据中解包一个记录"""
record = {}
offset = 0
for field_dict in fd_field_dicts:
slen = get_storage_length(field_dict)
record[field_dict['name']] = unpack_value(data[offset:offset + slen], field_dict)
offset += slen
return record
def write_variable_file(file_path: str, fd_field_dicts: list[dict],
records: list[dict]) -> int:
"""写入 RECORDING MODE V 文件(带 4 字节 RDW 前缀)。
RDW: Little-Endian unsigned short 记录长度(含自身4字节)
Args:
file_path: 输出路径
fd_field_dicts: FD 字段定义列表
records: 记录列表
Returns:
int: 写入的记录数
"""
with open(file_path, 'wb') as f:
for record in records:
data = bytearray()
for field_dict in fd_field_dicts:
val = record.get(field_dict['name'], '')
packed = pack_value(val, field_dict)
data.extend(packed)
record_len = len(data) + 4
rdw = struct.pack('<H', record_len)
f.write(rdw)
f.write(b'\x00\x00')
f.write(data)
logger.info(f" wrote {len(records)} records to {file_path}")
return len(records)