Files
cobol-java-v3/cobol_testgen/to_sql.py
T
hangshuo652 0203ead96b feat: config-driven multi-scenario + gcda accumulation fix + gcov merge fix
Config-driven architecture:
- coverage_dates: YAML-driven date replacement for LEAVE_RECORDS
- row_overrides: per-scenario field overrides (e.g. STATUS='9')
- delete_all_rows: per-scenario table emptying for empty-cursor tests

gcda accumulation fix (V3 bug #7):
- Delete .gcda before each scenario run to prevent GnuCOBOL accumulation
- Force-copy scenario gcda (skip mtime check)
- Copy scenario DB to CWD data/kin.db for CONNECT TO path resolution

Multi-run gcov merge fix:
- Always merge multi-run gcov data regardless of generate_coverage flag
- Fix v3_root UnboundLocalError in cleanup path

Other fixes:
- _make_key_unique: skip WHERE-constrained columns to avoid PK conflict
- incremental_supplement: support fields_dict for base record generation
- check_coverage: use structure coverage data if available
- orchestrator.py: filter _-prefixed fields in TestCase; merge Agent2Data cases
2026-07-18 08:42:55 +08:00

484 lines
16 KiB
Python

"""SQL层:WHERE约束解析 + DB输入行生成"""
import re
import logging
import itertools
logger = logging.getLogger(__name__)
# ── String literal protection ──
def _protect_strings(text: str) -> (str, list):
"""Replace string literals with placeholders. Returns (clean_text, replacements)."""
replacements = []
def _repl(m):
idx = len(replacements)
replacements.append(m.group(0))
return f"__STR{idx}__"
cleaned = re.sub(r"'[^']*'|\"[^\"]*\"", _repl, text)
return cleaned, replacements
def _restore_strings(text: str, replacements: list) -> str:
for i, s in enumerate(replacements):
text = text.replace(f"__STR{i}__", s)
return text
# ── Bracket-aware AND splitting ──
def _split_on_AND(text: str) -> list[str]:
"""Split WHERE clause on AND, respecting parentheses."""
parts = []
current = []
depth = 0
tokens = re.split(r'(\bAND\b|\bOR\b|[()])', text, flags=re.IGNORECASE)
for token in tokens:
if not token.strip():
continue
if token == '(':
depth += 1
current.append(token)
elif token == ')':
depth -= 1
current.append(token)
elif token.upper() == 'AND' and depth == 0:
parts.append(' '.join(current).strip())
current = []
elif token.upper() == 'OR' and depth == 0:
current.append(token) # OR stays as inner condition text
else:
current.append(token)
if current:
parts.append(' '.join(current).strip())
return parts
# ── WHERE condition parsing ──
_COL_OP_PAT = re.compile(
r'(\w[\w.-]*)\s*' # column name (with optional alias prefix)
r'(=|>|<|>=|<=|<>|!=|NOT\s*=)\s*'
r'(:\w[\w-]*(?::\w[\w-]*)?|__STR\d+__|[\w\d.-]+)',
re.IGNORECASE
)
_RE_IN_CLAUSE = re.compile(
r'(\w[\w.-]*)\s+(NOT\s+)?IN\s*\((.+?)\)',
re.IGNORECASE
)
_RE_BETWEEN = re.compile(
r'(\w[\w.-]*)\s+(NOT\s+)?BETWEEN\s+(.+?)\s+AND\s+(.+)',
re.IGNORECASE
)
_RE_LIKE = re.compile(
r'(\w[\w.-]*)\s+(NOT\s+)?LIKE\s+(__STR\d+__)',
re.IGNORECASE
)
_RE_IS_NULL = re.compile(
r'(\w[\w.-]*)\s+IS\s+(NOT\s+)?NULL',
re.IGNORECASE
)
def _parse_where_condition(part: str, replacements: list) -> dict | None:
"""Parse a single WHERE condition (after AND split)."""
part = part.strip()
if not part:
return None
# IS NULL
m = _RE_IS_NULL.match(part)
if m:
col = m.group(1).upper()
neg = bool(m.group(2))
return {'col': col, 'type': 'is_null', 'neg': neg, 'op': 'IS NULL' if not neg else 'IS NOT NULL'}
# IN
m = _RE_IN_CLAUSE.match(part)
if m:
col = m.group(1).upper()
neg = bool(m.group(2))
vals_text = m.group(3)
# Parse values from IN list
vals = []
for v in re.split(r'\s*,\s*', vals_text):
v = v.strip()
if v.startswith('__STR') and v.endswith('__'):
idx = int(v[5:-2])
vals.append(replacements[idx].strip("'\""))
elif v.startswith(':'):
vals.append({'type': 'host_var', 'host_var': v[1:].upper()})
else:
vals.append(v.strip())
return {
'col': col, 'type': 'in', 'neg': neg,
'op': 'NOT IN' if neg else 'IN',
'values': vals,
}
# BETWEEN
m = _RE_BETWEEN.match(part)
if m:
col = m.group(1).upper()
neg = bool(m.group(2))
lo = m.group(3).strip()
hi = m.group(4).strip()
return {
'col': col, 'type': 'between', 'neg': neg,
'op': 'BETWEEN',
'lo': lo.strip("'\""), 'hi': hi.strip("'\""),
}
# LIKE
m = _RE_LIKE.match(part)
if m:
col = m.group(1).upper()
neg = bool(m.group(2))
pat_ph = m.group(3)
idx = int(pat_ph[5:-2])
pattern = replacements[idx].strip("'\"") if idx < len(replacements) else pat_ph
return {
'col': col, 'type': 'like', 'neg': neg,
'op': 'NOT LIKE' if neg else 'LIKE',
'pattern': pattern,
}
# col op value
m = _COL_OP_PAT.match(part)
if m:
col = m.group(1).upper()
op = m.group(2).upper().strip()
val = m.group(3).strip()
# Normalize NOT = to <>
if op == 'NOT =' or op == 'NOT=':
op = '<>'
if val.startswith(':'):
host_var = val[1:].upper()
if ':' in host_var:
host_var = host_var.split(':')[0]
return {'col': col, 'type': 'host_var', 'host_var': host_var, 'op': op, 'literal': None}
elif val.startswith('__STR') and val.endswith('__'):
idx = int(val[5:-2])
_quotes = "'\""
literal = replacements[idx].strip(_quotes) if idx < len(replacements) else val
return {'col': col, 'type': 'literal', 'host_var': None, 'op': op, 'literal': literal}
else:
return {'col': col, 'type': 'literal', 'host_var': None, 'op': op, 'literal': val}
return None
# ── Column name → COBOL field name ──
_COLUMN_MAP = {}
def guess_cobol_field(col_name: str, table: str,
declared_columns: dict,
column_map: dict = None) -> str:
"""Map SQL column name to COBOL field name.
Priority: 1. DECLARE TABLE PIC alias 2. column_map 3. naming conv 4. as-is
"""
if column_map is None:
column_map = _COLUMN_MAP
# 1. DECLARE TABLE explicit PIC mapping
if table in declared_columns:
for c in declared_columns[table]:
if c['name'] == col_name and c.get('db_type') == 'PIC':
return c.get('pic', col_name)
# 2. User map
key = f"{table}.{col_name}"
if key in column_map:
return column_map[key]
# 3. Naming convention: CUST_ID → CUST-ID
candidate = col_name.replace('_', '-')
# 4. Strip table alias prefix: A.ID → ID
if '.' in candidate:
candidate = candidate.split('.')[1]
return candidate
# ── Main constraint extraction ──
def sql_extract_constraints(where_clause: str, table: str,
host_vars: dict[str, str],
column_map: dict[str, str],
declared_columns: dict) -> list[dict]:
"""Parse WHERE clause into constraint list."""
if not where_clause:
return []
# Protect string literals
cleaned, replacements = _protect_strings(where_clause)
# Split on AND
and_parts = _split_on_AND(cleaned)
constraints = []
for part in and_parts:
part = part.strip()
if not part:
continue
cond = _parse_where_condition(part, replacements)
if cond:
# Map column to COBOL field
cobol_field = guess_cobol_field(cond['col'], table, declared_columns, column_map)
cond['cobol_field'] = cobol_field
constraints.append(cond)
else:
logger.warning(f"Unparseable WHERE condition: {_restore_strings(part, replacements)}")
return constraints
# ── DB input row generation ──
_COLUMN_DEFAULTS = {
'CHAR': lambda size: ' ' * (size or 1),
'VARCHAR': lambda size: ' ' * (size or 1),
'INTEGER': lambda _: '000000000',
'SMALLINT': lambda _: '0000',
'DECIMAL': lambda _: '000000',
'DATE': lambda _: '20260603',
'PIC': lambda _: '?',
}
def _format_db_value(col_info: dict, raw_val: str) -> str:
db_type = col_info.get('db_type', 'CHAR')
formatter = _COLUMN_DEFAULTS.get(db_type, lambda _: str(raw_val)[:10])
default = formatter(0)
if raw_val is None:
return default
if db_type in ('INTEGER', 'SMALLINT', 'DECIMAL'):
try:
return str(int(raw_val)).zfill(len(default))
except ValueError:
return default
return str(raw_val).ljust(len(default))[:len(default)]
def _make_key_unique(key_val: str, path_index: int, seen_keys: set) -> str:
unique = f"{path_index:03d}{key_val[:5]}"
while unique in seen_keys:
unique = f"{path_index:03d}{hash(key_val) % 100000:05d}"
seen_keys.add(unique)
return unique
def collect_sql_meta(assignments: dict, declared_columns: dict,
column_map: dict = None) -> list[dict]:
"""Collect SQL metadata from assignments. Returns list of SQL info dicts."""
sql_meta = []
seen = set()
for tgt, asgn_list in assignments.items():
if isinstance(asgn_list, dict):
asgn_list = [asgn_list]
for asgn in asgn_list:
atype = asgn.get('type', '')
if not atype.startswith('exec_sql_'):
continue
if atype == 'exec_sql_fetch':
continue
key = asgn.get('sql_text', '')
if key in seen:
continue
seen.add(key)
where = asgn.get('where', '')
table = asgn.get('table', '')
where_constraints = sql_extract_constraints(
where, table, {}, column_map or {}, declared_columns
)
meta = dict(asgn)
meta['where_constraints'] = where_constraints
sql_meta.append(meta)
return sql_meta
def _path_has_sql_ok(path_cons: list) -> bool:
"""Check if a path requires SQLCODE = 0 (SQL succeeded)."""
sql_ok = True # default: no SQLCODE constraint, assume success
for pc in path_cons:
if len(pc) >= 4 and pc[0] == 'SQLCODE':
if pc[1] == '<>' and pc[3]:
sql_ok = False
if pc[1] == '=' and not pc[3]:
sql_ok = False
if pc[1] == '>' and pc[3]:
sql_ok = False
break
return sql_ok
def _infer_columns_from_where(where_cons: list) -> list[dict]:
"""Infer column definitions from WHERE constraints when DECLARE TABLE is missing."""
seen = {}
for wc in where_cons:
col_name = wc.get('col', '').split('.')[-1]
if col_name and col_name not in seen:
seen[col_name] = {'name': col_name, 'db_type': 'CHAR', 'size': 10}
return list(seen.values())
def _parse_select_columns(select_list: str) -> list[str]:
"""Parse SELECT column list into individual column names.
Handles: 'COL1, COL2' → ['COL1', 'COL2']
'COL1, COL2 AS alias' → ['COL1', 'COL2']
"""
if not select_list:
return []
cols = []
for part in select_list.split(','):
part = part.strip()
# Strip table alias prefix (T.COL → COL)
if '.' in part:
part = part.split('.')[1] if '.' in part else part
# Strip AS alias
m = re.search(r'^(\w[\w.-]*)', part)
if m:
cols.append(m.group(1).upper())
return cols
def build_db_input(
branch_paths: list[tuple[list, dict]],
fields_dict: list[dict],
assignments: dict,
sql_meta: list[dict],
declared_columns: dict,
records: list[dict] = None,
) -> dict:
"""Generate DB input rows per branch path.
Returns {table: [{col: val, ...}, ...]}.
"""
if not sql_meta:
return {}
db_input = {}
seen_keys = {}
seq_counter = itertools.count(1)
# Collect all SQL meta per path
for path_idx, (path_cons, path_assign) in enumerate(branch_paths):
# Skip paths where SQL fails (SQLCODE <> 0)
if not _path_has_sql_ok(path_cons):
continue
rec = records[path_idx] if records and path_idx < len(records) else {}
for sql in sql_meta:
atype = sql.get('type', '')
table = sql.get('table', '')
if not table:
continue
where_cons = sql.get('where_constraints', [])
if table not in db_input:
db_input[table] = []
seen_keys[table] = set()
if atype == 'exec_sql_insert':
# INSERT creates rows at runtime; no initial rows needed
continue
if atype in ('exec_sql_delete', 'exec_sql_update'):
# DELETE/UPDATE needs existing rows to act on
col_infos = declared_columns.get(table, [])
if not col_infos:
col_infos = _infer_columns_from_where(where_cons)
row = {}
for ci in col_infos:
col_name = ci['name'].upper()
val = None
for wc in where_cons:
wc_col = wc.get('col', '').upper().split('.')[-1]
if wc_col != col_name:
continue
if wc['type'] == 'literal':
val = wc.get('literal', '')
break
elif wc['type'] == 'host_var':
hv = wc.get('host_var', '').upper()
val = str(rec.get(hv, ''))
break
if val is None or not val.strip():
val = str(rec.get(ci['name'], ''))
if val and val.strip():
row[ci['name']] = _format_db_value(ci, val)
if not row:
row['_path'] = str(path_idx)
db_input[table].append(row)
continue
# exec_sql_select (and any future read-only types)
row = {}
col_infos = declared_columns.get(table, [])
if not col_infos:
col_infos = _infer_columns_from_where(where_cons)
into_vars = sql.get('into_vars', [])
# Map INTO vars → actual SQL column names from SELECT clause
select_cols = _parse_select_columns(sql.get('select_list', ''))
into_to_col: dict[str, str] = {}
for i, iv in enumerate(into_vars):
if i < len(select_cols):
into_to_col[iv] = select_cols[i]
# Use SQL column names (not INTO var names) for row keys
for sc in select_cols:
if sc not in [c['name'] for c in col_infos]:
col_infos.append({'name': sc, 'db_type': 'CHAR', 'size': 20})
where_cols = set()
for col_info in col_infos:
col_name = col_info['name']
val = None
for wc in where_cons:
if wc['type'] == 'literal' and wc.get('col', '').upper() == col_name:
val = wc.get('literal', '')
where_cols.add(col_name)
break
if wc['type'] == 'host_var':
hv = wc.get('host_var', '').upper()
for pc_field, pc_op, pc_val, pc_want in path_cons:
if pc_field == hv:
val = pc_val if pc_want else ''
break
if val is None and hv in rec:
val = str(rec[hv])
# Try to find value from INTO variable in the record
if val is None:
for iv, scola in into_to_col.items():
if scola == col_name and iv in rec:
val = str(rec[iv])
break
# Try COBOL field name mapping
if val is None:
cobol_field = guess_cobol_field(col_name, table, declared_columns)
if cobol_field in rec:
val = str(rec[cobol_field])
if val is not None:
row[col_name] = _format_db_value(col_info, val)
else:
row[col_name] = _format_db_value(col_info, str(next(seq_counter)))
if not row:
row['_path'] = str(path_idx)
if col_infos:
first_col = col_infos[0]['name']
if first_col in row and (not where_cols or first_col not in where_cols):
row[first_col] = _make_key_unique(row[first_col], path_idx, seen_keys[table])
db_input[table].append(row)
return db_input