276 lines
9.0 KiB
Python
276 lines
9.0 KiB
Python
"""COBOL 文件 I/O:DISPLAY/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()
|
||
# 本 GnuCOBOL(GC32-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)
|